Skip to content

[WS1][CUDA][Qwen3-Next] TP4 GDN, convolution and D=256 attention mixers on real weights (RFC #428 stack 6/8) - #507

Open
fusheng-ji wants to merge 71 commits into
RL-Align:test-qwennextfrom
fusheng-ji:feat/ws1-qwen3-next-tp4-mixers
Open

fusheng-ji wants to merge 71 commits into
RL-Align:test-qwennextfrom
fusheng-ji:feat/ws1-qwen3-next-tp4-mixers

Conversation

@fusheng-ji

@fusheng-ji fusheng-ji commented Oct 10, 2026 •

Copy link
Copy Markdown

Repository refactor update

Merged RL-Align/RL-Kernel:test-qwennext at a5a92ecf6cbf29fdcfc862846e53db5fa244c35e as requested in RFC #428. This PR still targets test-qwennext.

Implementation changes now live in the canonical rl_engine/backends, rl_engine/reference, rl_engine/runtime, and rl_engine/validation directories. Where present, model assembly and vLLM training bridges use rl_engine/models/qwen3_next and rl_engine/integrations/engines/train/vllm, respectively. Tests use tests/models/qwen3_next; evidence and validation commands use tools/validation/models and tools/validation/operators. Existing upstream compatibility entry points and pinned operator identifiers are preserved. CI commands and path filters follow the new layout, including the test-qwennext target.

Validation after migration: 366 passed in the relevant regression suites. CUDA tests used extensions built from the matching migrated native sources on B200. Syntax, undefined-name/duplicate-import checks, formatting, and git diff --check passed.

Additional checks: layout/packaging compatibility 32 passed, 1 skipped; workload/registry/validation compatibility 76 passed; both norms passed the canonical C3 forward and C4 gradient invariance gates. The provider suite passed 113 tests with one GPU; its mandatory cross-device test also passed when rerun with both B200 GPUs available (114 passed across the two runs). The four-GPU TP4 real-weight gates below were then rerun on the migrated stack (see Results).

Stack: #466 (merged) → #467 → #468 → #469 → #503 → this PR, base test-qwennext; merges after #503. GitHub shows the whole stack's commits; this PR's own diff is the commits on top of #503. Prep for RFC #428 full_model_chain; serves the gdn_state_bridge row (CUDA).

Status

  • Ready for review. No known acceptance failure.
  • The three TP4 Qwen3-Next mixers (GDN, full attention, MoE) pass their real-weight gates on 4 × B200, bitwise.
  • Known: the attention forward is slow at long prefill. RL-Kernel's deterministic attention materialises the FP32 probability matrix: 10 ms at 2048 tokens and 160 ms at 8192, against 0.21 / 0.31 ms for FlashAttention-4 (vendored by vLLM), which this PR's runner also measures as batch-invariant in the forward and the packed backward. This PR keeps RL-Kernel's kernel on both sides (train and rollout must run the same forward); moving to FlashAttention-4 is follow-up work, see the reuse decision below.

What

  • GDN (qwen3_next_gdn.py, qwen3_next_conv.py, qwen3_next_provider.py): packed recurrent prefill and decode through the pinned vLLM 0.30.0 provider, the depthwise causal convolution with its cache handoff, and one shared core whose prompt state crosses an explicit boundary into decode and into the trainer's gradient path.
  • TP4GDN (qwen3_next_tp_gdn.py): HF GQA-interleaved weight mapping per rank, explicit GDNState, Megatron-style TP boundaries from qwen3_next_tp.py.
  • Attention: csrc/cuda/attention/deterministic_attention.cu accepts head dim 256 on CUDA (ROCm stays 128) and validates backward inputs; shared_attention and TP4FullAttention with the KV head replicated on each rank pair, partial RoPE (64 of 256 dims) at absolute positions and per-slot KV state.
  • Gates: tools/validation/models/qwen3_next_tp_{gdn,attention,moe}_check.py on a shared tools/validation/models/qwen3_next_gate_common.py.

Prior art & reuse decision

GDN: reused as surveyed in #469 (the decode step runs vLLM's own FLA kernels; RL-Kernel supplies the state boundary and the golden). Attention, measured here (one TP4 rank: 4 Q heads, 1 KV head, D = 256, causal):

Implementation Batch BI Chunk (last 64) = prefill Decode (last 1) = prefill Backward, packed = alone Rel. error vs FP64
RL-Kernel deterministic attention yes yes yes dq / dk / dv 1.6e-3
FlashAttention-4 (CuTe DSL, vendored by vLLM 0.30.0) yes yes yes dq / dk / dv ¹ 2.0e-3
vLLM FA2, num_splits=1 yes yes yes no backward 2.0e-3
vLLM FA2, num_splits auto yes yes no no backward 2.0e-3
vLLM Triton unified attention (2D) yes yes no no backward 2.0e-3
Transformer Engine 2.16.1, inference (cuDNN 9.20 fused) yes yes yes no fused head-dim-256 backward on SM100 2.1e-3
Transformer Engine 2.16.1, training (unfused fallback) no error ² error ² no 3.6e-3
Megatron-core 0.16 local DotProductAttention yes ³ yes yes dq / dk / dv ³ 3.6e-3
torch SDPA 2.13 yes no no dq no 2.0e-3
FlashInfer 0.6.18.post1 ragged prefill no no no no backward 2.1e-3

¹ Default backward; its deterministic=True backward asserts for head dim 256 on SM100.
² TE's unfused THD path raises a shape error when the query is shorter than the KV.
³ No packed-sequence path: one launch per sequence, so batch invariance is by construction.

TP4 attention prior art

Latency (ms) RL-Kernel FA4 FA2 split=1 FA2 auto Triton 2D TE fused TE unfused Megatron SDPA FlashInfer
prefill 512 / 2048 / 8192 0.73 / 10.1 / 160 0.21 / 0.21 / 0.31 0.12 / 0.16 / 0.66 0.12 / 0.15 / 0.66 0.26 / 0.32 / 1.13 0.34 / 0.32 / 0.61 2.0 / 5.5 / 22.7 0.32 / 0.32 / 3.25 0.07 / 0.12 / 0.61 0.24 / 0.29 / 0.76
decode, 8192 KV 0.76 0.26 0.34 0.16 0.56 0.44 4.5 0.30 0.31 0.24
forward + backward, 2048 25.0 0.91 – – – – 10.3 1.34 0.74 –

Decision: keep RL-Kernel's deterministic attention on both sides in this PR. It is the only candidate here that is RL-Kernel-owned, covers forward and backward with one documented reduction order, and already backs the WS1 attention contract; train and rollout must run the same kernel. FlashAttention-4 is the strongest reuse candidate for the follow-up on the slow forward: batch-invariant forward and packed backward, and 30–500× faster at 2048–8192 tokens. Adopting it means pinning vLLM's vendored CuTe build on both sides and accepting a backward that has no deterministic mode at this head dim; both are out of scope here. TE fused is invariant for inference only; Megatron-core's local path only by launching per sequence.

Results (4 × B200, official checkpoint)

Gate Layer Cases (all bitwise, all 4 ranks)
GDN 0 HF load/export of the parameter storage; full == chunked prefill (output, conv and recurrent state) at 8/64/256/1024 tokens; response recompute == no-grad forward; gradient reaches the prompt through the state; replicated norm gradient — 7 / 7
attention 3 HF load/export with KV pair replicas; full == chunked prefill and KV continuation at 8–1024; variable-length batch of 4 == each alone, any order; 16-token decode == prefill rows; recompute; KV-pair and QK-norm gradient replicas — 10 / 10
MoE 0 as in #503, re-run at this commit — 12 / 12

Evidence (rank JSONs, attention report and figure, README with the exact commands) is in docs/usage/evidence/qwen3-next-tp4-mixers-b200/, measured from clean clones of 9e6c4f2 (gates) and 703c6dc (attention runner).

Rerun after the repository refactor: the three gates were rerun from a clean clone of the migrated stack tip (this PR plus #543 and #544 on top, f10a271) on 4 × B200 with the same checkpoint. All passed on all four ranks, and every recorded case value and digest in the 12 rank files is identical to the evidence above; only the MoE wall-clock timings differ.

Tests

python setup.py build_ext --inplace
python -m pytest -q tests/models/qwen3_next/test_qwen3_next_tp_gdn.py tests/models/qwen3_next/test_qwen3_next_tp_blocks.py tests/models/qwen3_next/test_qwen3_next_forward_contract.py
TRITON_F32_DEFAULT=ieee python -m pytest -q tests/models/qwen3_next/check_qwen3_next_{attention,gdn_bridge,gdn_sequence,conv_bridge,core_matrix,shared_core,forward}.py
torchrun --nproc-per-node 4 tools/validation/models/qwen3_next_tp_gdn_check.py --checkpoint <ckpt> --output <dir>
torchrun --nproc-per-node 4 tools/validation/models/qwen3_next_tp_attention_check.py --checkpoint <ckpt> --output <dir>
python tools/validation/models/qwen3_next_attention_prior_art.py --only <all but te_fused> --out <dir>/report-main.json
python tools/validation/models/qwen3_next_attention_prior_art.py --only te_fused --merge <dir>/report-main.json --out <dir>/report.json
Check Result
mixer GPU checks + CPU tests, one process 123 passed
existing attention suites (D = 128) 766 passed; 2 failures in test_attention_dispatch, also on main
real-weight gates 29 / 29 cases, 4 ranks; rerun after the refactor: all passed, identical results

The check_ files import vLLM and run in the Qwen3-Next-provider-GPU workflow; tests/models/qwen3_next/test_qwen3_next_tp_gdn.py joins ci.yml.

Scope

CUDA, TP4. ROCm keeps D = 128. No engine adapter or model-level claim; those are full_model_chain.

`test_registry_dispatches_rms_norm` asserts that the registry resolves
`rms_norm` to `RMSNormCudaOp` whenever CUDA and the compiled kernels are
both present, but `OpBackend` had no CUDA member for this operator and the
CUDA priority list contained only `PYTORCH_NATIVE_RMS_NORM`, so the assert
could never hold. The test therefore fails on any CUDA machine that builds
the native extension, and only passes when `_HAS_CUDA_RMSNORM` is false --
which is why an unbuilt CI has not caught it.

`RMSNormCudaOp` is already a first-class backend elsewhere: it is the
`"cuda"` candidate in `gtest/operator_specs.py` and is used directly by
`attention_preprocess.py`. Only the registry was missing it.

Add `OpBackend.CUDA_RMS_NORM` and put it ahead of the PyTorch reference in
the CUDA priority list. Because `_load_backend` only catches import errors
and this module imports cleanly without `_C`, a CUDA-first list would
otherwise hand out an op that raises at call time on an unbuilt install; so
`RMSNormCudaOp.__init__` now validates the extension and its three symbols,
matching `_require_cuda_activation` in the activation ops. The registry
already treats a backend whose construction raises as unavailable, so the
list degrades to `NativeRMSNormOp` as before.

Only the `cuda` priority map changes; rocm/musa/cpu/npu are untouched.

Verified on 2x B200 (sm_100, driver 580.126.20, torch 2.13.0+cu130) with the
extension rebuilt from source:

  pytest tests/test_rms_norm.py -q
  pytest tests/ rl_engine/tests/ -q -p no:randomly \
      --ignore=tests/test_rocm_aiter_api_contract.py

tests/test_rms_norm.py passes, including test_registry_dispatches_rms_norm, which
fails on the merge-base. The full suite gains no failure. The ignored file fails
to import on the merge-base as well.

Absolute suite counts are reported in the PR description against a named base
commit, not here: they shift whenever a sibling test is added, so a count frozen
in a commit message goes stale the moment the branch is rebased.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…back

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…ires

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
The forward, backward-dx and both backward-dw launchers took the current
CUDA stream without switching to the input's device, so a tensor on cuda:1
while cuda:0 is current launched on the wrong GPU. Add a device guard on
the input's device in each launcher. Use at::cuda::OptionalCUDAGuard with
the headers included unconditionally, as activation.cu does: those
launchers compile in the ROCm build too, where the file's existing
c10::cuda::CUDAGuard stays inside the !USE_ROCM block. A two-device test
checks that the op runs on the input's device and matches the
single-device result.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Advances RFC RL-Align#428 C1 on the CUDA track. Claim level: L0 repeatable and L1
batch-invariant. L2 is NOT claimed -- see "On exactness against vLLM" below.

Qwen3-Next's decoder and final norms store a zero-centred weight and compute
`x * rstd * (1 + w)`, with the `1 +` applied in fp32 after the upcast. Folding
it into a bf16 weight beforehand rounds the offset away, so it has to reach the
kernel as a parameter rather than being pre-applied by the caller.

  NativeRMSNormOp            gains `weight_offset` (default 0.0)
  Qwen3NextRMSNormOp         subclasses it, overriding only weight_offset = 1.0
  Qwen3NextRMSNormGatedOp    GDN gated norm: plain w, weight multiply in fp32
  Qwen3NextRMSNormGatedHFOp  the transformers convention, kept as a witness

The offset is applied under `if cls.weight_offset:` rather than unconditionally,
because `0.0 + w` rewrites -0.0 to +0.0. torch.equal does not notice that, but a
bitwise comparison does, and the plain path must stay bit-for-bit what it was.
A test pins it at the bit level.

The gated pair exists because transformers and vLLM disagree on where the gated
norm's weight multiply happens, and the gap is not a ULP: on bf16 /
head_v_dim=128 they differ in 35% of elements with max|diff| = 6.25e-2.
Isolating the cast order alone reproduces the gap (5.3e-2), so the cast order
dominates rather than the reduction order. The two conventions share their
validation and normalization and differ only in a `_scale_by_weight` hook.

CUDA: `weight_offset` added to the forward and dx kernels, defaulting to 0.0 so
every existing caller and binding is unaffected. The dw kernel is untouched:
d/dw (offset + w) == d/dw w.

  weight_offset=1.0 vs an explicit fp32 (1 + w) weight   bitwise equal
  weight_offset=1.0 vs a bf16-folded (1 + w) weight      differs, as required
  default offset vs the previous kernel                  bitwise equal

The first line is the correctness argument: the in-kernel offset is the same
arithmetic as the fp32 reference, not an approximation. The second is a
regression guard -- if it ever passes, the offset has stopped being fp32.

On exactness against vLLM
-------------------------
Measured over 40 seeds (bf16, head_v_dim=128, 512 rows):

  ours vs forward_native            6/40 seeds differ, worst 1.56e-2
  ours vs forward_cuda             18/40 seeds differ, worst 3.91e-3
  forward_native vs forward_cuda   21/40 seeds differ, worst 1.56e-2

vLLM's own two paths are not bitwise equal to each other, so "bitwise equal to
vLLM" is undefined until a single provider is named. In fp32 the two paths
differ on ~36% of elements, every one by an fp32 ULP -- the tree shapes differ,
the semantics do not. What this reproduces is the convention (fp32 weight
multiply, single trailing cast); the residual is the reduction tree.

The reduction stays the repo's fixed 32-wide chunked sum, which is what buys
L1. It was introduced for NPU but is needed on CUDA too: over 20 seeds at
H=2048 in bf16 -- Qwen3-Next's own hidden_size and dtype -- a plain mean(-1)
broke slice invariance on 1 of 20 while the chunked reduction broke on 0 of 20.
Matching stock vLLM bitwise would mean adopting a reduction that is not itself
batch-invariant, i.e. trading L1 for L2.

tests/check_qwen3_next_norm_providers.py pins the dispatch facts and bounds the
gap, asserting magnitudes rather than equality so a vLLM bump that changes the
provider fails loudly. It is named `check_` rather than `test_`, following
tests/distributed/check_*.py: it imports real vLLM, and
tests/test_framework_operator_integrations.py asserts vllm is absent from
sys.modules, an invariant any collected test importing vLLM would break for the
whole session.

Also: `rl_engine/_C.pyi` updated for the new `weight_offset` argument (CI runs
mypy against it), `tests/test_qwen3_next_norm.py` added to the CI test list in
.github/workflows/ci.yml, and an operator page added per
docs/operators/README.md ("the documentation page is part of the operator
contract").

Verified on 2x B200 (sm_100, driver 580.126.20, torch 2.13.0+cu130, triton
3.7.1, vllm 0.30.0, transformers 5.17.0):

  pytest tests/test_qwen3_next_norm.py -q
  pytest tests/check_qwen3_next_norm_providers.py -q
  pytest tests/ rl_engine/tests/ -q -p no:randomly \
      --ignore=tests/test_rocm_aiter_api_contract.py

Both new files pass and the full suite gains no failure. The ignored file fails
to import on the merge-base as well.

Absolute suite counts are reported in the PR description against a named base
commit, not here: they shift whenever a sibling test is added, so a count frozen
in a commit message goes stale the moment the branch is rebased.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
… VJPs

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Adds the CUDA kernel behind `Qwen3NextRMSNormGatedOp`, and registers both it and
the zero-centred decoder norm from the previous commit.

  y = x * rstd * (weight_offset + weight) * act(gate)

Every multiply is fp32 with one cast at the store, matching vLLM's RMSNormGated
with norm_before_gate=True and group_size=None -- the only configuration the GDN
block constructs. Other configurations are rejected rather than approximated.
`expf` is used rather than the `__expf` intrinsic: the fast intrinsic trades
accuracy for speed and would move the result away from the fp32 reference.

Forward reuses the existing block_reduce_sum / choose_threads(H), so for a fixed
H the reduction tree is independent of the row count -- that is the L1 guarantee
-- and `rstd` comes out bitwise identical to the ungated kernel for the same x,
which is asserted.

Backward is assembled from deterministic pieces:
  dx      new kernel, the ungated dx with (w + offset) -> (w + offset) * act(z)
  dweight reuses rmsnorm_dweight_rows_fp32 + the ascending-row fp32 left fold
  dgate   row-local and reduction-free, fp32 in the wrapper

Both backwards route through one `_fold_dweight_rows` helper so the file keeps a
single left-fold entrypoint, which tests/test_vjp_fp32.py pins, and one
`_require_cuda_symbols` helper so the module has a single availability contract.

The gated op is deliberately NOT a subclass of RMSNormCudaOp: it takes an extra
required tensor, so it cannot stand in for one. Same reasoning as on the PyTorch
side, and stated in its docstring so the question is not reopened.

Registration: `rms_norm_gated` and `qwen3_next_rms_norm` in OP_SPECS (both with a
cuda-sm90 candidate, as every other reduction spec carries), operator_inputs
builders, OpBackend members and priority maps on all five platforms,
test_dispatch assertions, and the WS1 registered-ops set in
tests/test_ws1_gtest_gpu.py. Operator pages added per docs/operators/README.md,
which states the page is part of the operator contract.

Claim level: L0 repeatable and L1 batch-invariant. NOT L2 -- see
tests/check_qwen3_next_norm_providers.py, which pins the dispatch facts and
bounds the gap against vLLM rather than asserting equality.

Verified on 2x B200 (sm_100, driver 580.126.20, torch 2.13.0+cu130, vllm 0.30.0):

  pytest tests/test_qwen3_next_norm.py -q
  pytest tests/check_qwen3_next_norm_providers.py -q
  python scripts/check_operator.py --op rms_norm_gated --candidate cuda \
      --device cuda --dtype bf16 --check-grad        -> pass_rate=1.0000
  python scripts/check_operator.py --op qwen3_next_rms_norm --candidate cuda \
      --device cuda --dtype bf16 --check-grad        -> pass_rate=1.0000
  pytest tests/ rl_engine/tests/ -q -p no:randomly \
      --ignore=tests/test_rocm_aiter_api_contract.py

Both new files pass, both operators report pass_rate=1.0000 on the CUDA
candidate, and the full suite gains no failure. The ignored file fails to import
on the merge-base as well.

Absolute suite counts are reported in the PR description against a named base
commit, not here: they shift whenever a sibling test is added, so a count frozen
in a commit message goes stale the moment the branch is rebased.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…t norms

Tests: the low-precision and CUDA-vs-golden comparisons hard-coded
atol=2e-2/rtol=1.6e-2, which is the contract's elementwise bf16 row. They now
resolve forward_accuracy for op_class="reduction" from tolerance_contract.json.
That loosens bf16 (5e-2/2e-2) and tightens fp16 (1e-3/1e-3, previously 2e-2);
the fp16 cases still pass. The vLLM provider-gap bounds in
tests/check_qwen3_next_norm_providers.py are labelled as gap bounds, not
contract thresholds.

Docs and docstrings: correct statements this branch had committed.
- The operator page no longer says the op is registered or prints a
  check_operator command; the gtest spec and registry entry arrive with the
  gated-norm PR.
- Withdrawn: "7 of 1048576 differ" (single seed, no script), the 35% / 5.3e-2
  cast-order isolation (an fp32 round-trip is a no-op), "needed on CUDA as
  well", and "matching vLLM means trading L1 for L2" (one unreproduced
  observation). Replaced with the scoped claim levels and the probe results.
- The provider check no longer claims to establish which path vLLM
  dispatches, or a per-element fp32 ULP bound.
- The module docstring is cut to the contract and links to the operator page.

Comments: why the kernel's offset add is guarded (signed zero, with the
pinning test), and why parameter_vjp_contributions_fp32 passes the offset.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…atistics

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…ion fallback

The zero-centred norm was registered CUDA-first alongside the gated one, but only
rms_norm_gated had dispatch tests. Mirror them: pin the priority on all five
platforms, and assert that the inherited RMSNormCudaOp.__init__ symbol check makes
the registry fall through to the PyTorch reference.

Each assertion was checked to fail on CPU with the behaviour removed: PyTorch
first in the CUDA list, the ROCm entry dropped, and the __init__ check removed.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…ctivation, offset

The bitwise rstd identity was asserted for one configuration only (bf16, H=128,
silu, offset 0). Parametrize it over fp32/fp16/bf16 x H in {128, 2048, 5120} x
{silu, sigmoid} x offset in {0.0, 1.0}, 36 cases, with the plain kernel taking
the same offset.

All 36 are expected to hold by construction: both kernels run the same
sum-of-squares loop, block_reduce_sum and choose_threads(H) launch shape, and
neither the gate nor the offset enters the statistic. The H values cover both
block sizes (128 and 512 threads) and two per-thread serial lengths. Not yet run
on a GPU; the cases skip without the compiled extension.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…16, 17)

R2-1 (private bf16 tolerances cited as the reduction row) was resolved in the
preceding merge, which switched the gated tests to c1's _forward_tol.

- R2-2 docs/operators/qwen3-next-rms-norm-gated.md: Accuracy now points to
  tolerance_contract.json (forward_accuracy, reduction x dtype); the 35% / 6.25e-2
  HF-vs-vLLM figure is labelled a one-off with its setup (one seed, B200, bf16,
  head_v_dim=128, 512 rows, eager); the 40-seed figures are attributed to
  tests/check_qwen3_next_norm_providers.py. Withdrawn, as on the c1 page: the
  cast-order isolation, the per-element fp32 ULP wording, and "fails closed (RFC
  section 6 item 7)" for configurations the op has no parameter for. Claim levels
  now read L1 (prefix slices; batch size, chunking, padding, permutation in the
  C3/C4 gates), L0 not separately tested; the rstd sweep is described.
- R2-7 csrc/ops.cpp: drop the dim and size TORCH_CHECKs that rmsnorm_check_weight
  already performs, in rmsnorm_forward, rmsnorm_backward_dx and
  rmsnorm_gated_check. Error messages for those cases now come from
  rmsnorm_check_weight; no test matches the removed messages.
- R2-8 workload_report() in rl_engine/testing/ws1_workload.py replaces the
  payload block duplicated in both gate scripts. full_model_evidence is read from
  the manifest: False for the Qwen3-Next norm manifest, None (not declared) for the
  Dense manifest, which previously reported a hard-coded False.
- R2-9 the gate scripts import validate_norm_dimensions at module level, and
  --manifest has a help string.
- R2-12 the four dispatch tests become two, parametrized over both norms. The
  fallback test now resolves for device="cuda", so the CUDA-first list is walked
  even on a CPU-only host, and asserts the CUDA backend was tried and failed.
  Before, on CPU, get_op went through the CPU list and the assertion held
  trivially.
- R2-13 SPDX headers on rl_engine/testing/qwen3_next_workload.py and
  tests/test_qwen3_next_workload.py.
- R2-14 ci/run_ws1_gtest.sh runs the four Qwen3-Next gates as explicit commands
  instead of a loop over interpolated script names.
- R2-16 Qwen3NextRMSNormGatedCudaOp takes activation as a constructor argument,
  validated before the extension check; the test no longer mutates the instance.
  "swish" stays as vLLM's alias for "silu", now asserted bitwise on CUDA.
- R2-17 the Python dgate path adds weight_offset only when it is nonzero, like the
  kernels, so a -0.0 weight keeps its sign in dgate.

New tests: activation rejected at construction (CPU), swish == silu (CUDA), dgate
signed zero (CUDA, 3 dtypes). Checked on CPU that the fallback and constructor
tests fail with the behaviour removed. The CUDA tests and the csrc change are not
yet run on a GPU.

Deferred: R2-6 (shared sum-of-squares device function), R2-10 (private validator
imports), R2-11.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
… guard

The six rmsnorm launchers compiled on both CUDA and ROCm (plain forward and
dx, gated forward and dx, partial and reduce dweight) constructed a
non-optional c10::cuda::CUDAGuard. No other file the ROCm build compiles uses
that class outside a USE_ROCM guard, while activation.cu and
deterministic_attention.cu use at::cuda::OptionalCUDAGuard. Match the guard used
by activation.cu so the ROCm build relies only on patterns it already compiles.

The guard follows the same tensor as before (x, or partial_dw for the reduce
launcher). CUDA behaviour is unchanged: these tensors are always on a CUDA
device, so the optional guard always sets the device. The left-fold launcher
sits inside the USE_ROCM guard and keeps its c10::cuda::CUDAGuard. The ROCm
build is still untested.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…th the stacked C1 branch

Carried over from the former merge of the C1 branch into this one: the decoder-norm page again states the registration and the check_operator command, the gated tests use the contract's reduction tolerances through _forward_tol, and the reference docstrings name the gated page. Same final tree as before the re-stack.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
The trainer-side reference for a Qwen3-Next rollout decode step: the recurrent
Gated DeltaNet update and the causal-conv1d state update that precedes it.
Claim level: L0 repeatable and L1 batch-invariant; L2 is not claimed.

Design notes, measurements and the deferral list are in
docs/design/ws1-c6-428-gdn-recurrent-replay.md.

Provider
--------
`fused_recurrent_gated_delta_rule_packed_decode`, not
`fused_sigmoid_gating_delta_rule_update`. A decode-only, non-speculative batch
returns early at qwen_gdn_linear_attn.py:1295-1307 into the packed path because
VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE defaults to true; the sigmoid-gating
kernel is the fallback for mixed prefill+decode and spec decode. The env
defaults that decide this are asserted, so a vLLM bump fails loudly.

Three places the kernel and the HF model disagree, transcribed from the kernel:

  * The gating is fused -- beta = sigmoid(b) and
    g = -exp(A_log) * softplus(a + dt_bias) are computed in fp32 inside the
    kernel with a threshold branch at 20, not as separate PyTorch ops.
  * No repeat_interleave: the kernel indexes i_h = i_hv // (HV // H).
  * The QK norm is an L2 norm over a plain sum, x / sqrt(sum(x*x) + 1e-6),
    dividing by sqrt rather than multiplying by rsqrt. scale hits q after the
    norm; k is never scaled.

A fourth candidate divergence turned out not to be one: prefill passes
use_qk_l2norm_in_kernel=False only because fused_post_conv_prep(apply_l2norm=True)
already normalized q/k. Both paths normalize exactly once.

Conv accumulation order is load-bearing: taps accumulate sequentially,
acc = acc + win[t] * w[t] from zero. With an fp32 cache that reproduces the
provider bitwise; a tree sum over the same four terms does not, and neither
does an FMA.

State ABI is mirrored: paged [num_blocks, HV, V, K] and
[num_blocks, dim, width-1], NULL_BLOCK_ID (<= 0) skipping, an fp32 accumulator,
and a store that rounds to the cache dtype -- fp32 or bf16, per
FUSED_GDN_STATE_DTYPES. Contractions run in the repo's fixed 32-wide chunk order
via one `_chunked_sum` primitive, which is what the L1 claim rests on.

Agreement, B200, Qwen3-Next dims (H=16, HV=32, K=V=128)
-------------------------------------------------------
  recurrent, fp32 state, B=1..64   max|d out| 1.5e-08..6.1e-05, state <= 3.0e-07
  recurrent, bf16 state, B=1..64   max|d out| 3.7e-09..3.1e-05, state <= 2.0e-03
  conv: the rolled state is BITWISE exact in every configuration
  conv, fp32 cache: output bitwise but for 15 of 524288 elements at B=64
  conv, bf16 cache: output agrees on ~63%, each disagreement one bf16 ULP
  L1: a sequence's output and state block are bitwise identical alone or at any
      position in a batch of 64, for both state dtypes

Where the provider rounds in the bf16-conv-cache case is not reproduced;
recorded as an open gap rather than guessed at.

Decode versus chunked prefill
-----------------------------
  T=8    max|d| 3.66e-04 fp32 state / 5.49e-04 bf16   (7.0e-03 / 1.0e-02 rel)
  T=64   4.88e-04 / 5.49e-04                          (6.6e-03 / 7.5e-03 rel)
  T=256  3.66e-04 / 3.66e-04                          (5.0e-03 / 5.0e-03 rel)

~0.5-1% relative and flat in T. The recurrence is contracting -- per-step decay
exp(g) averages ~0.47 -- so old rounding error is forgotten at roughly the rate
old signal is, and the fp32-vs-bf16 state drift saturates at ~2% relative on the
state (2.67e-03 at step 1, 1.52e-02 at 64, 2.47e-02 at 1024: a 16x longer run
past step 64 grows it 1.6x).

So the two paths are not bitwise, and ~1% relative on logits is still material
for RL importance ratios, but this is a bounded error rather than a divergence.
An fp32 recurrent state remains the right choice; the reason is the plateau, not
a blow-up. Separately, the chunked prefill kernel refuses fp32 q/k/v outright
(chunk.py:213), and causal_conv1d_update does not bounds-check
conv_state_indices under its default validate_data=False.

tests/check_gdn_recurrent_golden.py is named check_ rather than test_, following
tests/distributed/check_*.py: it imports real vLLM, and
tests/test_framework_operator_integrations.py asserts vllm is absent from
sys.modules.

  pytest tests/check_gdn_recurrent_golden.py -q
  pytest tests/ rl_engine/tests/ -q -p no:randomly \
      --ignore=tests/test_rocm_aiter_api_contract.py

The new file passes and the full suite gains no failure.

Absolute suite counts are reported in the PR description against a named base
commit, not here: they shift whenever a sibling test is added, so a count frozen
in a commit message goes stale the moment the branch is rebased.

Not covered, deliberately: speculative decode and MTP, a backward for the
recurrence, and the provider bridge. Reasons in the design doc.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…touched

test_null_block_id_is_skipped_by_both checked that both sides write zeros
for a NULL_BLOCK_ID row, but asserted "block untouched" only for the golden.
The provider returns before any store for state_idx <= 0
(fused_recurrent.py:299-303 in vllm 0.30.0); assert it on the provider's
state as well so the test name matches what it checks.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
ws1-gtest-gpu.yml: 158fc75 added four Qwen3-Next/GDN test files to the
pull_request and push path filters together with a step in
ci/run_ws1_gtest.sh that ran the two check_ files. 03176a3 moved that step
to qwen3-next-provider-gpu.yml but left the filters, so editing those files
triggered a RunPod job that never executed them. Remove the filters; the
file is back to its state on feat/cuda-qwen3-next-gated-rmsnorm.

qwen3-next-provider-gpu.yml: drop rl_engine/integrations/qwen3_next_gdn.py
and ci/run_qwen3_next_operator_gates.sh from the path filter. Neither file
exists on this branch.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
The design note quoted provider-vs-golden figures (max|d out|, max|d state|,
conv mismatch counts) that no committed code printed, so they could not be
reproduced. scripts/ws1_gdn_provider_agreement.py prints them as JSON,
using tests/check_gdn_recurrent_golden.py's own input helpers and seeds
(imported, not copied):

- recurrent: per (batch, state dtype), max|d out|, max|d state| and
  bitwise mismatch counts, next to the bounds the check file asserts;
- conv: per (batch, cache dtype), output mismatch count, max|diff| and
  whether the rolled state is bitwise equal;
- the conv comparison again with triton.knobs.language.default_fp_fusion
  off, plus the fma.rn.f32 count in each compiled variant of the
  provider's conv-update kernel, to test whether the fp32-cache mismatches
  are FMA contraction.

Provenance (git commit and dirty flag, torch/triton/vllm versions, device)
is recorded with the results. Requires CUDA and vLLM; not run in CI.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Review after 03176a3 found nine statements in the note that the source
does not support. Fix them:

- §1: the decode path is chosen per engine step and by the model, not by
  env defaults alone. Tabulate decode-only / mixed / draft-token steps with
  their conv and recurrent kernels. Qwen3-Next's gqa_interleaved_layout=True
  makes the fused CUDA decode unsupported, so the layer falls back to triton
  and fused_gdn_decode_post_conv_mtp is unreachable (:1834). The env-default
  assertion guards the packed-decode default but not VLLM_GDN_DECODE_KERNEL.
- §2: the golden follows the kernel's arithmetic but is not a transcription
  everywhere: softplus uses log1p where the kernel uses log(1 + exp), and
  FLA_USE_FAST_OPS swaps in fast exp/log.
- §3: split NULL_BLOCK_ID semantics. The recurrent provider skips <= 0 and
  writes zeros; the conv provider skips only == null_block_id (0) and does
  not write the output row; both goldens skip <= 0 and write zeros.
- §4: replace the "measured" table, which had no runner, with the bounds the
  check file asserts and point to scripts/ws1_gdn_provider_agreement.py.
  State the fp32-cache conv mismatch and the FMA hypothesis the runner tests.
- §5: add the mixed decode-and-prefill caveat on the provider side, including
  vllm-project/vllm#49827, and state that validate_data=True does not check
  index values either.
- §6: restate the MTP deferral against RFC RL-Align#428 §2.2 and the actual
  Qwen3-Next MTP path.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
_git() returned the string "unavailable: ..." on failure, so
bool(_git("status", "--porcelain")) turned "git missing / not a checkout"
into git_dirty=True. Return None on failure and record git_dirty and
git_commit as null in that case, so the provenance distinguishes unknown
from dirty.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
The §6 row deferred backward because "a naive BPTT through a sequential
recurrence is not batch-invariant". RFC RL-Align#428 does not ask backward to be:
§2.2 item 4 says backward need not match rollout, only be correct for the
replayed forward, and §9.1 makes the GDN backward its own work item, C7.
Cite those instead, matching the PR description.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
A run of the runner at d66e595 showed both of its FMA probes were ineffective:

- The fusion-off arm reused the fused kernel. Triton 3.7.1 keys its
  in-memory kernel cache on the launch kwargs; the default_fp_fusion knob
  is read in parse_options only after a cache miss. Flipping the knob
  in-process never missed, so all six recorded variants had
  enable_fp_fusion=True and the "off" counts were the "on" counts again.
  Run the fusion-off arm in a child process with TRITON_DEFAULT_FP_FUSION=0
  and a private TRITON_CACHE_DIR instead.
- fma.rn.f32 = 0 in the PTX does not rule out FMA. With fusion on, Triton
  emits plain mul.f32/add.f32 and lets ptxas contract them into FFMA; with
  fusion off it passes --fmad=false. Count plain and .rn PTX ops, and
  FFMA/FMUL/FADD in the SASS (Triton's bundled cuobjdump).

Report per-arm kernel variants under conv_kernels and add fusion_check,
which says whether the fusion-off variants were actually compiled without
fusion, so a silent no-op cannot pass for a result again.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
The runner on d66e595 (B200, clean checkout) reproduced the recurrent
figures the note had withdrawn for lack of a runner, so restore them as a
per-batch table with their provenance and bitwise output-mismatch counts.

The conv counts quoted in test_conv_output_matches_provider_with_fp32_cache's
docstring (1 of 8192 at B=1, 15 of 524288 at B=64) predate the product-
rounding fix; the runner measures 0, 0, 1 and 5 for B = 1, 4, 17, 64.
Replace them, and record the bf16-cache counts in the note.

State plainly that the FMA explanation for the fp32-cache mismatches is not
determined: that run's fusion-off arm reused the fused kernel, and a PTX
without fma.rn.f32 cannot exclude ptxas contraction. The runner now tests
both (previous commit).

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…tches

The fixed runner (b54e50f) on 7622616, B200, settles the FMA hypothesis
left open in §4. Its fusion-off arm took effect: all six variants compiled
with enable_fp_fusion=False, and their PTX uses only .rn-qualified f32
mul/add, which ptxas may not contract. With contraction verifiably off,
the fp32-cache mismatch counts (0, 0, 1, 5 for B = 1, 4, 17, 64) and
max|diff| are identical to the fused run, and neither arm's SASS has an
FFMA. So FP contraction is ruled out; the actual cause stays open.

The note bases the conclusion on the A/B outcome, not on the opcode
counts: each variant's PTX has only two f32 multiplies, too few to be the
four tap products, so the counters do not show where the products are
computed.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
§5 said a mixed decode-and-prefill step leaves the decode row's bf16 output
equal. That holds only at TP1 head counts (H=16, HV=32). At TP4 per-rank
head counts (H=4, HV=8) the same step already moves one output element
(4.8e-05 relative) and 32,663 of 131,072 state elements (1.7e-07), with the
same numbers for all three prefill-bearing compositions.

Tabulate both head counts and state the measurement's limits: one
standalone layer from config.json, synthetic parameters and cache, single
process with no engine, scheduler or CUDA graph, metadata built by the test.
The 4-of-20 propagation figure is TP1-only and does not carry over to real
weights. Conv output and state matched bitwise, so the difference is in the
recurrent kernels. vllm-project/vllm#49827's two commits on 0.30.0 close the
decode-row gap at both head counts (scheduler part untested); with packed
decode disabled the decode row is bitwise equal in every composition.

Source: the GDN evidence branch's committed results (not yet published).
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…rouped expert GEMM

RFC RL-Align#428 moe_route_combine_contract (CUDA).

- Routing: FP32 router through the pinned vLLM batch-invariant GEMM, a
  fixed-order softmax, stable descending argsort (ties go to the lower
  expert id) and FP32 renormalisation of the ten selected weights.
- Experts: per-route gate_up and down projections through vLLM 0.30.0's
  fused_moe Triton kernel with the batch-invariant tile passed explicitly
  (BLOCK 64/64/32, SPLIT_K=1); no routed-weight multiply and no sum inside
  the kernel, so every (token, route) row has one writer. Bitwise equal to
  the previous per-expert GEMM loop at TP4 and TP1 widths; fails closed if
  VLLM_TRITON_USE_TD is enabled.
- Combine: ten ordered FP32 multiply-adds and one BF16 cast; no atomics.
- Backward: per-expert pinned GEMMs into one dense gradient buffer; dx is
  summed in ascending expert order.
- TP4 block: HF load/export for all 512 experts, replicated router and
  shared-expert gate, Megatron-style copy/reduce autograd boundaries moved
  into qwen3_next_tp.py.
- Prior-art runner and plot comparing HF transformers, vLLM fused_moe
  (BI=0/1) and FlashInfer cutlass_fused_moe; TP4 real-weight gate script.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Prep for RFC RL-Align#428 full_model_chain; serves the gdn_state_bridge row.

- GDN: packed recurrent prefill/decode through the pinned vLLM 0.30.0
  provider with an explicit state boundary (qwen3_next_gdn.py), the
  depthwise causal convolution with cache handoff (qwen3_next_conv.py),
  their combined shared core (qwen3_next_provider.py) and the TP4 block
  with HF GQA-interleaved weight mapping (qwen3_next_tp_gdn.py), now on the
  shared TP boundaries in qwen3_next_tp.py.
- Attention: the deterministic CUDA attention kernel accepts head dim 256
  (ROCm stays 128) and validates backward inputs; shared_attention and
  TP4FullAttention with the KV head replicated on each rank pair, partial
  RoPE at absolute positions and explicit per-slot KV state.
- Real-weight TP4 gates for GDN (layer 0) and attention (layer 3), sharing
  scripts/qwen3_next_gate_common.py with the MoE gate.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…nce (B200, 2e950f3)

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Compares RL-Kernel's deterministic attention with torch SDPA, vLLM FlashAttention-2 (num_splits auto and 1), vLLM's Triton unified attention and FlashInfer's ragged prefill: batch, chunked-prefill and decode invariance, backward repeatability, accuracy against FP64 and latency.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…nce (B200, 9e6c4f2/6517077)

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…r-art runner

- Megatron-core 0.16 MoELayer with Transformer Engine grouped GEMM and fused
  permute, configured as VIME runs Qwen3-Next (FP32 softmax router, top-10,
  all-to-all dispatcher), including its backward.
- SGLang 0.5.21 Triton fused_moe in default and deterministic-inference mode.
  sgl_kernel cannot load next to torch 2.13, so its two kernels on this path
  are replaced by SGLang's own Triton moe_sum_reduce and JIT
  moe_align_block_size; any other sgl_kernel call fails loudly.
- --only and --merge combine reports from different Python environments.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…ured (B200, c58816b)

Every library in the reuse table is now measured: Megatron-core 0.16 + TE 2.16 as VIME configures Qwen3-Next, and SGLang 0.5.21 in default and deterministic-inference mode. The TP4 gate evidence is unchanged (2e950f3, identical MoE code).

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…gatron-core attention

Add three engines to the attention prior-art runner:
- FlashAttention-4 (CuTe DSL) varlen as vendored by vLLM, forward and backward.
- Transformer Engine DotProductAttention on THD-packed sequences, in
  inference mode (cuDNN fused) and in training mode (on SM100 TE has no
  fused head-dim-256 backward and falls back to its unfused path).
- Megatron-core's local DotProductAttention, one launch per sequence.

The backward check now compares the target's dq/dk/dv alone against the
same sequence packed first and last, so engines with a packed backward are
checked for batch invariance rather than repeatability. A prefill/decode
crash is recorded per candidate without skipping its other checks.
--only and --merge combine reports from different Python environments.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
The decode and forward+backward latency sections are keyed by candidate
directly, the prefill section by size first; merge them at any depth.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…egatron-core measured (B200, 703c6dc)

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
@coderabbitai

coderabbitai Bot commented Oct 10, 2026 •

Copy link
Copy Markdown

Important

Review skipped

Auto reviews are disabled on base/target branches other than the default branch.

Please check the settings in the CodeRabbit UI or the .coderabbit.yaml file in this repository. To trigger a single review, invoke the @coderabbitai review command.

⚙️ Run configuration
  • Configuration used: defaults
  • Review profile: CHILL
  • Plan: Advanced
  • Run ID: 7f51aefe-741c-4ed5-a480-965ec522517c

You can disable this status message by setting the reviews.review_status to false in the CodeRabbit configuration file.

Use the checkbox below for a quick retry:

  • 🔍 Trigger review
  • Autofix · Keep fixing CodeRabbit findings and required CI, and resolving merge conflicts

Comment @coderabbitai help to get the list of available commands.

…p as E402-exempt

The gate scripts insert the repo root into sys.path before importing
rl_engine and tools, which flake8 reports as E402. Mark those imports
with noqa, as the other tools/ scripts do.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
…p as E402-exempt

The gate scripts insert the repo root into sys.path before importing
rl_engine and tools, which flake8 reports as E402. Mark those imports
with noqa, as the other tools/ scripts do.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Annotate the latency samples and accept the default-argument lambdas used as
timed calls. No behaviour change.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>
… runner

Annotate the latency samples and accept the default-argument lambdas used as
timed calls. No behaviour change.

Signed-off-by: Wenbo Ji <36562829+fusheng-ji@users.noreply.github.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant